UnsortedSegmentSum

对输入张量按给定的 segment 索引进行无序分段求和操作。

对于每个输入元素,根据其对应的 index 值,将其累加到输出张量中 对应的 segment 位置。不同 segment 之间的顺序不做任何排序保证。

\[\text{output}_{s, j} = \sum_{i \,|\, index_i = s} \text{input}_{i, j}\]

其中:

  • \(i \in [0, \text{dim0})\)

  • \(j \in [0, \text{dim1})\)

  • \(s \in [0, \text{id\_max})\)

输入:
  • input - 输入张量的数据地址。 数据类型需与所调用的 UnsortedSegmentSum 接口类型一致。

  • index - segment 索引数组地址,类型为 int*, 长度为 dim0,用于指定每一行输入数据所属的 segment。

  • dim0 - 输入张量的第 0 维大小,即输入元素的行数。

  • dim1 - 输入张量的第 1 维大小,即每行输入数据的特征维度。

  • id_max - segment 的最大数量,决定输出张量的第 0 维大小。

  • core_mask - 核掩码(仅共享存储版本使用)。

输出:
  • output - 输出张量的数据地址, 形状为 [id_max, dim1], 数据类型与 input 保持一致。

支持平台:

FT78NE MT7004

备注

  • 该算子为无序分段求和(Unsorted),不保证 segment 内或 segment 间的顺序。

  • FT78NE 支持int8, int16, int32, fp32, fp64, cplx64, cplx128

  • MT7004 支持fp16, fp32, int16, int32, cplx64

  • 输出张量在计算前会被初始化为 0。

  • index 中的取值范围应满足 0 <= index[i] < id_max

共享存储版本:

void i8_unsorted_segment_sum_s(int8_t *input, int *index, int8_t *output, int dim0, int dim1, int id_max, int core_mask)
void i16_unsorted_segment_sum_s(int16_t *input, int *index, int16_t *output, int dim0, int dim1, int id_max, int core_mask)
void i32_unsorted_segment_sum_s(int32_t *input, int *index, int32_t *output, int dim0, int dim1, int id_max, int core_mask)
void hp_unsorted_segment_sum_s(float16 *input, int *index, float16 *output, int dim0, int dim1, int id_max, int core_mask)
void fp_unsorted_segment_sum_s(float *input, int *index, float *output, int dim0, int dim1, int id_max, int core_mask)
void dp_unsorted_segment_sum_s(double *input, int *index, double *output, int dim0, int dim1, int id_max, int core_mask)
void c64_unsorted_segment_sum_s(float *input, int *index, float *output, int dim0, int dim1, int id_max, int core_mask)
void c128_unsorted_segment_sum_s(double *input, int *index, double *output, int dim0, int dim1, int id_max, int core_mask)

C调用示例:

 1// MT7004 示例(共享存储多核,DDR 地址)
 2#include <stdio.h>
 3#include <stdlib.h>
 4#include <string.h>
 5#include "mt7004/utils.h"
 6
 7int main(int argc, char* argv[]) {
 8    float *input = (float *)0x81000000;
 9    int *index = (int *)0x82000000;
10    float *output = (float *)0x83000000;
11
12    int i;
13    int dim0 = 128;
14    int dim1 = 64;
15    int id_max = 16;
16    int core_mask = 0xff;
17
18    srand(12345);
19    for (i = 0; i < dim0 * dim1; ++i) {
20        input[i] = (float)(rand() % 20 - 10);
21    }
22    for (i = 0; i < dim0; ++i) {
23        index[i] = rand() % id_max;
24    }
25    memset(output, 0, id_max * dim1 * sizeof(float));
26
27    fp_unsorted_segment_sum_s(input, index, output, dim0, dim1, id_max, core_mask);
28    return 0;
29}

私有存储版本:

void i8_unsorted_segment_sum_p(int8_t *input, int *index, int8_t *output, int dim0, int dim1, int id_max)
void i16_unsorted_segment_sum_p(int16_t *input, int *index, int16_t *output, int dim0, int dim1, int id_max)
void i32_unsorted_segment_sum_p(int32_t *input, int *index, int32_t *output, int dim0, int dim1, int id_max)
void hp_unsorted_segment_sum_p(float16 *input, int *index, float16 *output, int dim0, int dim1, int id_max)
void fp_unsorted_segment_sum_p(float *input, int *index, float *output, int dim0, int dim1, int id_max)
void dp_unsorted_segment_sum_p(double *input, int *index, double *output, int dim0, int dim1, int id_max)
void c64_unsorted_segment_sum_p(float *input, int *index, float *output, int dim0, int dim1, int id_max)
void c128_unsorted_segment_sum_p(double *input, int *index, double *output, int dim0, int dim1, int id_max)

C调用示例:

 1// MT7004 示例(私有存储单核,AM 地址)
 2#include <stdio.h>
 3#include <stdlib.h>
 4#include "mt7004/utils.h"
 5
 6int main(int argc, char* argv[]) {
 7    float *input = (float *)0x10010000;
 8    int *index = (int *)0x10020000;
 9    float *output = (float *)0x10030000;
10
11    int i;
12    int dim0 = 64;
13    int dim1 = 32;
14    int id_max = 8;
15
16    srand(12345);
17    for (i = 0; i < dim0 * dim1; ++i) {
18        input[i] = (float)(rand() % 20 - 10);
19    }
20    for (i = 0; i < dim0; ++i) {
21        index[i] = rand() % id_max;
22    }
23
24    fp_unsorted_segment_sum_p(input, index, output, dim0, dim1, id_max);
25    return 0;
26}